Skip to content

[AMD] gfx950 assembly attention: length-aware split-KV for dynamic workload - #39172

Merged
HaiShaw merged 2 commits into
sgl-project:mainfrom
zijiecode:pr/vattn-segplan
Sep 14, 2026
Merged

HaiShaw merged 2 commits into
sgl-project:mainfrom
zijiecode:pr/vattn-segplan

Conversation

@zijiecode

@zijiecode zijiecode commented Sep 12, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

#37465 splits every sequence of a verify / draft batch into the same number of KV segments, max(1, min(64, pow2floor(256 / (bs * kv_heads)))), one work-group per (segment, sequence, kv head). That is the right split when all sequences have similar length, not optimal when one batch holds requests that varies heavily in length. Two things then happen:

  • the layer finishes when the longest request's segments finish, so short requests' work-groups sit idle while a 248k request streams 8x longer segments;
  • above 16 sequences the per-sequence segment count halves (bs 17-32 -> 8 segments; on TP4 with 2 kv heads per rank bs 9-16 -> 8, bs 17-32 -> 4), so a 248k request runs as 31k-token segments and the kernel degrades to Triton speed. Half of the decode steps of the TP2 c32 leaderboard point are at bs 17-30.

This PR picks the segment length from the batch's actual sequence lengths and hands the kernel a compact work list, so every sequence gets as many segments as its length needs while the total work-group count stays at one per CU.

Modifications

The planned split is the default for bs > 1; bs == 1 keeps the fixed split of #37465 (plan_ptr == 0), which already uses 64 segments there. No new switch. Kernel build and the gfx950 / ROCm gating are unchanged from #37465.

  • vattn_asm_gfx950/__init__.py: a small Triton plan kernel (grid = num_seqs, graph-capture safe) picks one segment length T from seq_lens (smallest T with total segments <= num_CUs / kv_heads and per-sequence segments <= seg_max) and writes a work list plan[1 + i] = seq << 16 | seg (-1 past the end) plus a per-query-token segment count. The attention kernel is launched as a 1-D grid over the work list; bs == 1 bypasses the plan. The plan is built once per forward (cached on the seq_lens / cu_seqlens_q tensors, cleared at the start of every forward, keyed on the stream-capturing flag).
  • vattn3_core.s: kernarg 128 -> 136 B (plan_ptr); the prologue decodes (seq, seg) from the work list and exits on -1 before touching memory; tiles per segment come from plan[0].
  • vred.s: kernarg 56 -> 64 B (tok_nseg_ptr); the reduce reads only that token's segment count.
  • unified_attention_3d_mtp.py / aiter_backend.py: reset_verify_attn_plan_cache() called from init_forward_metadata and init_forward_metadata_out_graph.

Accuracy Tests

  • test/registered/kernel/attention/test_vattn_segplan.py (stage-b-test-1-gpu-small-amd-mi35x, 22 s on MI355X): 16 cases (GQA 16 and 8; uniform 16x70k; agentic skew including 0/1/5/17-token sequences; bs 1 / 2 / 24 / 64; ragged q 1..4) against an fp32 torch reference and against the fixed split. Same error vs. reference as the fixed split (max 0.0141 bf16 on the skew case, dominated by the fp8 KV); plan vs. fixed split differ by at most 0.007 (different fp32 reduction order across segment boundaries). Plan contents (T, work list, tok_nseg) checked against a Python reference in every case.
  • same file, plan cache: 15 calls in one forward -> 1 plan launch; reset -> recompute; in-place seq_lens change or a different tensor -> recompute; cached output bit-identical to uncached.
  • GSM8K, 5-shot, all 1319 questions, temperature 0, max new tokens 2048:
Attention kernel Accuracy
baseline 0.921
this PR 0.924

Benchmarking and Profiling

shape Triton unified_attention_3d_mtp asm, fixed split (#37465) asm, this PR
verify 16 x 119k uniform 379 us (2.6 TB/s) 176 us (5.5) 176 us (5.5)
verify bs16, agentic length distribution, 1.85M tokens 518 us (1.8) 266 us (3.6) 179 us (5.3)
verify bs20, agentic length distribution, 2.32M tokens 516 us (2.3) 482 us (2.5) 254 us (4.7)
draft extend ragged q=1..4, bs20 agentic 748 us (1.6) 472 us (2.5) 251 us (4.7)
bs1 110k / 248k - 44 / 79 us same (bypass)

End to end, InferenceX agentic recipe (Qwen3.5-397B-A17B MXFP4, MTP, aiperf trace replay; official harness mode: full aiperf, 3600 s, 10 warmups), MI355X, same tree and flags; off = fixed split of #37465, on = this PR. TP4 pairs ran one server at a time on the same GPUs; TP2 pairs ran off and on simultaneously on separate GPU pairs of the same node.

point TPOT p90 (ms) off on interactivity p90 (tok/s/user) off on throughput / GPU
TP2 c32 HiCache 28.02 27.00 (-3.6%) 35.7 37.0 (+3.6%) +0.9%
TP2 c8 6.12 5.85 (-4.4%) 163.5 170.9 (+4.5%) +0.5%
TP4 c16 7.19 6.86 (-4.6%) 139.1 145.8 (+4.8%) +0.5%
TP4 c4 4.31 4.11 (-4.6%) 231.9 243.5 (+5.0%) 0.0%

Server-side accept length is identical between off and on at every point (3.38-3.40). Same-tree A/B via a development-only switch that this PR does not ship. The p90 tail improves by 4-5% at every point (a long request joining a small batch no longer runs as 8 or fewer oversized segments); the median moves only at TP2 c32 (TPOT p50 12.47 -> 11.57 ms, -7.2%), where half of the decode steps are at bs 17-30 and the fixed split degrades. A same-configuration repeat on this replay differs by about 0.8%.

Checklist

  • Format your code according to the Format code with pre-commit guide.
  • Add unit tests (test/registered/kernel/attention/test_vattn_segplan.py).
  • Update documentation.
  • Provide accuracy and speed benchmark results.
  • Follow the SGLang code style guidance.

CI States

Latest PR Test (Base): ✅ Run #34686493077
Latest PR Test (Extra): ❌ Run #34686492967
Latest PR Test (AMD ROCm 10): ❌ Run #34686493122

@zijiecode zijiecode changed the title [AMD] gfx950 assembly attention: length-aware split-KV for agentic batches [AMD] gfx950 assembly attention: length-aware split-KV for dynamic workload Sep 12, 2026
…entic batches

For bs > 1 a small Triton plan kernel picks the segment length from the batch's
sequence lengths and writes a compact work list; the asm attention kernel runs as
a 1-D grid over that list (idle work-groups exit before touching memory) and the
reduce reads a per-token segment count. The plan is built once per forward and
cached on the seq_lens / cu_seqlens_q tensors (keyed on the stream-capturing
flag). bs == 1 keeps the fixed split of sgl-project#37465 (plan_ptr == 0), which already
uses 64 segments there.
@yichiche yichiche added the run-ci CI: run the baseline test suite on this PR label Sep 12, 2026

@yichiche yichiche left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Split-K optimization for the MTP assembly kernel; changes are AMD-only. LGTM.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

jit-kernel run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants